Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T05:02:23.783Z Public web read
NIP-34 coordinate30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omega
MaintainersHidden in public view
References2 branches · 1 tag
Read-only clonegit clone https://openagents.com/git/tenant.openagents/omega.git
Browse files

http.rs

809 lines · 31.3 KB · rust
1use anyhow::{Result, anyhow};
2use async_trait::async_trait;
3use collections::HashMap;
4use futures::{Stream, StreamExt};
5use gpui::BackgroundExecutor;
6use http_client::{AsyncBody, HttpClient, Request, Response, http::Method};
7use parking_lot::Mutex as SyncMutex;
8use std::{pin::Pin, sync::Arc};
9
10use crate::oauth::{self, OAuthTokenProvider, WwwAuthenticate};
11use crate::transport::Transport;
12use crate::types;
13
14/// Typed errors returned by the HTTP transport that callers can downcast from
15/// `anyhow::Error` to handle specific failure modes.
16#[derive(Debug)]
17pub enum TransportError {
18    /// The server returned 401 and token refresh either wasn't possible or
19    /// failed. The caller should initiate the OAuth authorization flow.
20    AuthRequired { www_authenticate: WwwAuthenticate },
21}
22
23impl std::fmt::Display for TransportError {
24    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25        match self {
26            TransportError::AuthRequired { .. } => {
27                write!(f, "OAuth authorization required")
28            }
29        }
30    }
31}
32
33impl std::error::Error for TransportError {}
34
35// Constants from MCP spec
36const HEADER_SESSION_ID: &str = "Mcp-Session-Id";
37const HEADER_PROTOCOL_VERSION: &str = "MCP-Protocol-Version";
38const EVENT_STREAM_MIME_TYPE: &str = "text/event-stream";
39const JSON_MIME_TYPE: &str = "application/json";
40
41/// HTTP Transport with session management and SSE support
42pub struct HttpTransport {
43    http_client: Arc<dyn HttpClient>,
44    endpoint: String,
45    session_id: Arc<SyncMutex<Option<String>>>,
46    /// Negotiated MCP protocol version, populated by `set_protocol_version`
47    /// after the initialize handshake. From 2025-06-18 onward the server
48    /// requires clients to echo this in the `MCP-Protocol-Version` header on
49    /// every subsequent request.
50    protocol_version: Arc<SyncMutex<Option<String>>>,
51    executor: BackgroundExecutor,
52    response_tx: async_channel::Sender<String>,
53    response_rx: async_channel::Receiver<String>,
54    error_tx: async_channel::Sender<String>,
55    error_rx: async_channel::Receiver<String>,
56    /// Static headers to include in every request (e.g. from server config).
57    headers: HashMap<String, String>,
58    /// When set, the transport attaches `Authorization: Bearer` headers and
59    /// handles 401 responses with token refresh + retry.
60    token_provider: Option<Arc<dyn OAuthTokenProvider>>,
61    /// The challenge from the last 401 this transport gave up on; cleared at
62    /// the start of each send so it always describes the most recent attempt.
63    /// See [`Transport::auth_challenge`].
64    auth_challenge: SyncMutex<Option<WwwAuthenticate>>,
65}
66
67impl HttpTransport {
68    pub fn new(
69        http_client: Arc<dyn HttpClient>,
70        endpoint: String,
71        headers: HashMap<String, String>,
72        executor: BackgroundExecutor,
73    ) -> Self {
74        Self::new_with_token_provider(http_client, endpoint, headers, executor, None)
75    }
76
77    pub fn new_with_token_provider(
78        http_client: Arc<dyn HttpClient>,
79        endpoint: String,
80        headers: HashMap<String, String>,
81        executor: BackgroundExecutor,
82        token_provider: Option<Arc<dyn OAuthTokenProvider>>,
83    ) -> Self {
84        let (response_tx, response_rx) = async_channel::unbounded();
85        let (error_tx, error_rx) = async_channel::unbounded();
86
87        Self {
88            http_client,
89            executor,
90            endpoint,
91            session_id: Arc::new(SyncMutex::new(None)),
92            protocol_version: Arc::new(SyncMutex::new(None)),
93            response_tx,
94            response_rx,
95            error_tx,
96            error_rx,
97            headers,
98            token_provider,
99            auth_challenge: SyncMutex::new(None),
100        }
101    }
102
103    /// Build a POST request for the given message body, attaching all standard
104    /// headers (content-type, accept, session ID, static headers, and bearer
105    /// token if available).
106    fn build_request(&self, message: &[u8]) -> Result<http_client::Request<AsyncBody>> {
107        let mut request_builder = Request::builder()
108            .method(Method::POST)
109            .uri(&self.endpoint)
110            .header("Content-Type", JSON_MIME_TYPE)
111            .header(
112                "Accept",
113                format!("{}, {}", JSON_MIME_TYPE, EVENT_STREAM_MIME_TYPE),
114            );
115
116        for (key, value) in &self.headers {
117            request_builder = request_builder.header(key.as_str(), value.as_str());
118        }
119
120        // Attach bearer token when a token provider is present.
121        if let Some(token) = self.token_provider.as_ref().and_then(|p| p.access_token()) {
122            request_builder = request_builder.header("Authorization", format!("Bearer {}", token));
123        }
124
125        // Add session ID if we have one (except for initialize).
126        if let Some(ref session_id) = *self.session_id.lock() {
127            request_builder = request_builder.header(HEADER_SESSION_ID, session_id.as_str());
128        }
129
130        // Echo the negotiated protocol version once initialization has
131        // completed. Required by servers speaking MCP 2025-06-18 or later.
132        if let Some(ref version) = *self.protocol_version.lock()
133            && types::requires_protocol_version_header(version)
134        {
135            request_builder = request_builder.header(HEADER_PROTOCOL_VERSION, version.as_str());
136        }
137
138        Ok(request_builder.body(AsyncBody::from(message.to_vec()))?)
139    }
140
141    /// Record the challenge so it remains observable after the failed send
142    /// tears down the client (see [`Transport::auth_challenge`]), and build
143    /// the typed error for the send itself.
144    fn auth_required(&self, www_authenticate: WwwAuthenticate) -> anyhow::Error {
145        *self.auth_challenge.lock() = Some(www_authenticate.clone());
146        TransportError::AuthRequired { www_authenticate }.into()
147    }
148
149    /// Send a message and handle the response based on content type.
150    async fn send_message(&self, message: String) -> Result<()> {
151        // The same server instance can be restarted over this transport; a
152        // challenge recorded by a previous client generation must not be
153        // observed by the current one.
154        *self.auth_challenge.lock() = None;
155
156        let is_notification =
157            !message.contains("\"id\":") || message.contains("notifications/initialized");
158
159        // If we currently have no access token, try refreshing before sending
160        // the request so restored but expired sessions do not need an initial
161        // 401 round-trip before they can recover.
162        if let Some(ref provider) = self.token_provider {
163            if provider.access_token().is_none() {
164                provider.try_refresh().await.unwrap_or(false);
165            }
166        }
167
168        let request = self.build_request(message.as_bytes())?;
169        let mut response = self.http_client.send(request).await?;
170
171        // On 401, try refreshing the token and retry once.
172        if response.status().as_u16() == 401 {
173            let www_auth_header = response
174                .headers()
175                .get("www-authenticate")
176                .and_then(|v| v.to_str().ok())
177                .unwrap_or("Bearer");
178
179            let www_authenticate =
180                oauth::parse_www_authenticate(www_auth_header).unwrap_or(WwwAuthenticate {
181                    resource_metadata: None,
182                    scope: None,
183                    error: None,
184                    error_description: None,
185                });
186
187            if let Some(ref provider) = self.token_provider {
188                if provider.try_refresh().await.unwrap_or(false) {
189                    // Retry with the refreshed token.
190                    let retry_request = self.build_request(message.as_bytes())?;
191                    response = self.http_client.send(retry_request).await?;
192
193                    // If still 401 after refresh, give up.
194                    if response.status().as_u16() == 401 {
195                        return Err(self.auth_required(www_authenticate));
196                    }
197                } else {
198                    return Err(self.auth_required(www_authenticate));
199                }
200            } else {
201                return Err(self.auth_required(www_authenticate));
202            }
203        }
204
205        // Handle different response types based on status and content-type.
206        match response.status() {
207            status if status.is_success() => {
208                // Check content type
209                let content_type = response
210                    .headers()
211                    .get("content-type")
212                    .and_then(|v| v.to_str().ok());
213
214                // Extract session ID from response headers if present
215                if let Some(session_id) = response
216                    .headers()
217                    .get(HEADER_SESSION_ID)
218                    .and_then(|v| v.to_str().ok())
219                {
220                    *self.session_id.lock() = Some(session_id.to_string());
221                    log::debug!("Session ID set: {}", session_id);
222                }
223
224                match content_type {
225                    Some(ct) if ct.starts_with(JSON_MIME_TYPE) => {
226                        // JSON response - read and forward immediately
227                        let mut body = String::new();
228                        futures::AsyncReadExt::read_to_string(response.body_mut(), &mut body)
229                            .await?;
230
231                        // Only send non-empty responses
232                        if !body.is_empty() {
233                            self.response_tx
234                                .send(body)
235                                .await
236                                .map_err(|_| anyhow!("Failed to send JSON response"))?;
237                        }
238                    }
239                    Some(ct) if ct.starts_with(EVENT_STREAM_MIME_TYPE) => {
240                        // SSE stream - set up streaming
241                        self.setup_sse_stream(response).await?;
242                    }
243                    _ => {
244                        // For notifications, 202 Accepted with no content type is ok
245                        if is_notification && status.as_u16() == 202 {
246                            log::debug!("Notification accepted");
247                        } else {
248                            return Err(anyhow!("Unexpected content type: {:?}", content_type));
249                        }
250                    }
251                }
252            }
253            status if status.as_u16() == 202 => {
254                // Accepted - notification acknowledged, no response needed
255                log::debug!("Notification accepted");
256            }
257            _ => {
258                let mut error_body = String::new();
259                futures::AsyncReadExt::read_to_string(response.body_mut(), &mut error_body).await?;
260
261                self.error_tx
262                    .send(format!("HTTP {}: {}", response.status(), error_body))
263                    .await
264                    .map_err(|_| anyhow!("Failed to send error"))?;
265            }
266        }
267
268        Ok(())
269    }
270
271    /// Set up SSE streaming from the response
272    async fn setup_sse_stream(&self, mut response: Response<AsyncBody>) -> Result<()> {
273        let response_tx = self.response_tx.clone();
274        let error_tx = self.error_tx.clone();
275
276        // Spawn a task to handle the SSE stream
277        self.executor
278            .spawn(async move {
279                let reader = futures::io::BufReader::new(response.body_mut());
280                let mut lines = futures::AsyncBufReadExt::lines(reader);
281
282                let mut data_buffer = Vec::new();
283                let mut in_message = false;
284
285                while let Some(line_result) = lines.next().await {
286                    match line_result {
287                        Ok(line) => {
288                            if line.is_empty() {
289                                // Empty line signals end of event
290                                if !data_buffer.is_empty() {
291                                    let message = data_buffer.join("\n");
292
293                                    // Filter out ping messages and empty data
294                                    if !message.trim().is_empty() && message != "ping" {
295                                        if let Err(e) = response_tx.send(message).await {
296                                            log::error!("Failed to send SSE message: {}", e);
297                                            break;
298                                        }
299                                    }
300                                    data_buffer.clear();
301                                }
302                                in_message = false;
303                            } else if let Some(data) = line.strip_prefix("data: ") {
304                                // Handle data lines
305                                let data = data.trim();
306                                if !data.is_empty() {
307                                    // Check if this is a ping message
308                                    if data == "ping" {
309                                        log::trace!("Received SSE ping");
310                                        continue;
311                                    }
312                                    data_buffer.push(data.to_string());
313                                    in_message = true;
314                                }
315                            } else if line.starts_with("event:")
316                                || line.starts_with("id:")
317                                || line.starts_with("retry:")
318                            {
319                                // Ignore other SSE fields
320                                continue;
321                            } else if in_message {
322                                // Continuation of data
323                                data_buffer.push(line);
324                            }
325                        }
326                        Err(e) => {
327                            let _ = error_tx.send(format!("SSE stream error: {}", e)).await;
328                            break;
329                        }
330                    }
331                }
332            })
333            .detach();
334
335        Ok(())
336    }
337}
338
339#[async_trait]
340impl Transport for HttpTransport {
341    async fn send(&self, message: String) -> Result<()> {
342        self.send_message(message).await
343    }
344
345    fn receive(&self) -> Pin<Box<dyn Stream<Item = String> + Send>> {
346        Box::pin(self.response_rx.clone())
347    }
348
349    fn receive_err(&self) -> Pin<Box<dyn Stream<Item = String> + Send>> {
350        Box::pin(self.error_rx.clone())
351    }
352
353    fn set_protocol_version(&self, version: &str) {
354        *self.protocol_version.lock() = Some(version.to_string());
355    }
356
357    fn auth_challenge(&self) -> Option<WwwAuthenticate> {
358        self.auth_challenge.lock().clone()
359    }
360}
361
362impl Drop for HttpTransport {
363    fn drop(&mut self) {
364        // Try to cleanup session on drop
365        let http_client = self.http_client.clone();
366        let endpoint = self.endpoint.clone();
367        let session_id = self.session_id.lock().clone();
368        let protocol_version = self.protocol_version.lock().clone();
369        let headers = self.headers.clone();
370        let access_token = self.token_provider.as_ref().and_then(|p| p.access_token());
371
372        if let Some(session_id) = session_id {
373            self.executor
374                .spawn(async move {
375                    let mut request_builder = Request::builder()
376                        .method(Method::DELETE)
377                        .uri(&endpoint)
378                        .header(HEADER_SESSION_ID, &session_id);
379
380                    // Add static authentication headers.
381                    for (key, value) in headers {
382                        request_builder = request_builder.header(key.as_str(), value.as_str());
383                    }
384
385                    // Attach bearer token if available.
386                    if let Some(token) = access_token {
387                        request_builder =
388                            request_builder.header("Authorization", format!("Bearer {}", token));
389                    }
390
391                    // Stamp the negotiated MCP protocol version on the DELETE
392                    // too, matching what `build_request` does for POSTs.
393                    if let Some(ref version) = protocol_version
394                        && types::requires_protocol_version_header(version)
395                    {
396                        request_builder =
397                            request_builder.header(HEADER_PROTOCOL_VERSION, version.as_str());
398                    }
399
400                    let request = request_builder.body(AsyncBody::empty());
401
402                    if let Ok(request) = request {
403                        let _ = http_client.send(request).await;
404                    }
405                })
406                .detach();
407        }
408    }
409}
410
411#[cfg(test)]
412mod tests {
413    use super::*;
414    use async_trait::async_trait;
415    use gpui::TestAppContext;
416    use parking_lot::Mutex as SyncMutex;
417    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
418
419    /// A mock token provider that returns a configurable token and tracks
420    /// refresh attempts.
421    struct FakeTokenProvider {
422        token: SyncMutex<Option<String>>,
423        refreshed_token: SyncMutex<Option<String>>,
424        refresh_succeeds: AtomicBool,
425        refresh_count: AtomicUsize,
426    }
427
428    impl FakeTokenProvider {
429        fn new(token: Option<&str>, refresh_succeeds: bool) -> Arc<Self> {
430            Self::with_refreshed_token(token, None, refresh_succeeds)
431        }
432
433        fn with_refreshed_token(
434            token: Option<&str>,
435            refreshed_token: Option<&str>,
436            refresh_succeeds: bool,
437        ) -> Arc<Self> {
438            Arc::new(Self {
439                token: SyncMutex::new(token.map(String::from)),
440                refreshed_token: SyncMutex::new(refreshed_token.map(String::from)),
441                refresh_succeeds: AtomicBool::new(refresh_succeeds),
442                refresh_count: AtomicUsize::new(0),
443            })
444        }
445
446        fn set_token(&self, token: &str) {
447            *self.token.lock() = Some(token.to_string());
448        }
449
450        fn refresh_count(&self) -> usize {
451            self.refresh_count.load(Ordering::SeqCst)
452        }
453    }
454
455    #[async_trait]
456    impl OAuthTokenProvider for FakeTokenProvider {
457        fn access_token(&self) -> Option<String> {
458            self.token.lock().clone()
459        }
460
461        async fn try_refresh(&self) -> Result<bool> {
462            self.refresh_count.fetch_add(1, Ordering::SeqCst);
463
464            let refresh_succeeds = self.refresh_succeeds.load(Ordering::SeqCst);
465            if refresh_succeeds {
466                if let Some(token) = self.refreshed_token.lock().clone() {
467                    *self.token.lock() = Some(token);
468                }
469            }
470
471            Ok(refresh_succeeds)
472        }
473    }
474
475    fn make_fake_http_client(
476        handler: impl Fn(
477            http_client::Request<AsyncBody>,
478        ) -> std::pin::Pin<
479            Box<dyn std::future::Future<Output = anyhow::Result<Response<AsyncBody>>> + Send>,
480        > + Send
481        + Sync
482        + 'static,
483    ) -> Arc<dyn HttpClient> {
484        http_client::FakeHttpClient::create(handler) as Arc<dyn HttpClient>
485    }
486
487    fn json_response(status: u16, body: &str) -> anyhow::Result<Response<AsyncBody>> {
488        Ok(Response::builder()
489            .status(status)
490            .header("Content-Type", "application/json")
491            .body(AsyncBody::from(body.as_bytes().to_vec()))
492            .unwrap())
493    }
494
495    #[gpui::test]
496    async fn test_bearer_token_attached_to_requests(cx: &mut TestAppContext) {
497        // Capture the Authorization header from the request.
498        let captured_auth = Arc::new(SyncMutex::new(None::<String>));
499        let captured_auth_clone = captured_auth.clone();
500
501        let client = make_fake_http_client(move |req| {
502            let auth = req
503                .headers()
504                .get("Authorization")
505                .map(|v| v.to_str().unwrap().to_string());
506            *captured_auth_clone.lock() = auth;
507            Box::pin(async { json_response(200, r#"{"jsonrpc":"2.0","id":1,"result":{}}"#) })
508        });
509
510        let provider = FakeTokenProvider::new(Some("test-access-token"), false);
511        let transport = HttpTransport::new_with_token_provider(
512            client,
513            "http://mcp.example.com/mcp".to_string(),
514            HashMap::default(),
515            cx.background_executor.clone(),
516            Some(provider),
517        );
518
519        transport
520            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
521            .await
522            .expect("send should succeed");
523
524        assert_eq!(
525            captured_auth.lock().as_deref(),
526            Some("Bearer test-access-token"),
527        );
528    }
529
530    #[gpui::test]
531    async fn test_no_bearer_token_without_provider(cx: &mut TestAppContext) {
532        let captured_auth = Arc::new(SyncMutex::new(None::<String>));
533        let captured_auth_clone = captured_auth.clone();
534
535        let client = make_fake_http_client(move |req| {
536            let auth = req
537                .headers()
538                .get("Authorization")
539                .map(|v| v.to_str().unwrap().to_string());
540            *captured_auth_clone.lock() = auth;
541            Box::pin(async { json_response(200, r#"{"jsonrpc":"2.0","id":1,"result":{}}"#) })
542        });
543
544        let transport = HttpTransport::new(
545            client,
546            "http://mcp.example.com/mcp".to_string(),
547            HashMap::default(),
548            cx.background_executor.clone(),
549        );
550
551        transport
552            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
553            .await
554            .expect("send should succeed");
555
556        assert!(captured_auth.lock().is_none());
557    }
558
559    #[gpui::test]
560    async fn test_missing_token_triggers_refresh_before_first_request(cx: &mut TestAppContext) {
561        let captured_auth = Arc::new(SyncMutex::new(None::<String>));
562        let captured_auth_clone = captured_auth.clone();
563
564        let client = make_fake_http_client(move |req| {
565            let auth = req
566                .headers()
567                .get("Authorization")
568                .map(|v| v.to_str().unwrap().to_string());
569            *captured_auth_clone.lock() = auth;
570            Box::pin(async { json_response(200, r#"{"jsonrpc":"2.0","id":1,"result":{}}"#) })
571        });
572
573        let provider = FakeTokenProvider::with_refreshed_token(None, Some("refreshed-token"), true);
574        let transport = HttpTransport::new_with_token_provider(
575            client,
576            "http://mcp.example.com/mcp".to_string(),
577            HashMap::default(),
578            cx.background_executor.clone(),
579            Some(provider.clone()),
580        );
581
582        transport
583            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
584            .await
585            .expect("send should succeed after proactive refresh");
586
587        assert_eq!(provider.refresh_count(), 1);
588        assert_eq!(
589            captured_auth.lock().as_deref(),
590            Some("Bearer refreshed-token"),
591        );
592    }
593
594    #[gpui::test]
595    async fn test_invalid_token_still_triggers_refresh_and_retry(cx: &mut TestAppContext) {
596        let request_count = Arc::new(AtomicUsize::new(0));
597        let request_count_clone = request_count.clone();
598
599        let client = make_fake_http_client(move |_req| {
600            let count = request_count_clone.fetch_add(1, Ordering::SeqCst);
601            Box::pin(async move {
602                if count == 0 {
603                    Ok(Response::builder()
604                        .status(401)
605                        .header(
606                            "WWW-Authenticate",
607                            r#"Bearer error="invalid_token", resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource""#,
608                        )
609                        .body(AsyncBody::from(b"Unauthorized".to_vec()))
610                        .unwrap())
611                } else {
612                    json_response(200, r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)
613                }
614            })
615        });
616
617        let provider = FakeTokenProvider::with_refreshed_token(
618            Some("old-token"),
619            Some("refreshed-token"),
620            true,
621        );
622        let transport = HttpTransport::new_with_token_provider(
623            client,
624            "http://mcp.example.com/mcp".to_string(),
625            HashMap::default(),
626            cx.background_executor.clone(),
627            Some(provider.clone()),
628        );
629
630        transport
631            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
632            .await
633            .expect("send should succeed after refresh");
634
635        assert_eq!(provider.refresh_count(), 1);
636        assert_eq!(request_count.load(Ordering::SeqCst), 2);
637    }
638
639    #[gpui::test]
640    async fn test_401_triggers_refresh_and_retry(cx: &mut TestAppContext) {
641        let request_count = Arc::new(AtomicUsize::new(0));
642        let request_count_clone = request_count.clone();
643
644        let client = make_fake_http_client(move |_req| {
645            let count = request_count_clone.fetch_add(1, Ordering::SeqCst);
646            Box::pin(async move {
647                if count == 0 {
648                    // First request: 401.
649                    Ok(Response::builder()
650                        .status(401)
651                        .header(
652                            "WWW-Authenticate",
653                            r#"Bearer resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource""#,
654                        )
655                        .body(AsyncBody::from(b"Unauthorized".to_vec()))
656                        .unwrap())
657                } else {
658                    // Retry after refresh: 200.
659                    json_response(200, r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)
660                }
661            })
662        });
663
664        let provider = FakeTokenProvider::new(Some("old-token"), true);
665        // Simulate the refresh updating the token.
666        let provider_ref = provider.clone();
667        let transport = HttpTransport::new_with_token_provider(
668            client,
669            "http://mcp.example.com/mcp".to_string(),
670            HashMap::default(),
671            cx.background_executor.clone(),
672            Some(provider.clone()),
673        );
674
675        // Set the new token that will be used on retry.
676        provider_ref.set_token("refreshed-token");
677
678        transport
679            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
680            .await
681            .expect("send should succeed after refresh");
682
683        assert_eq!(provider_ref.refresh_count(), 1);
684        assert_eq!(request_count.load(Ordering::SeqCst), 2);
685    }
686
687    #[gpui::test]
688    async fn test_401_returns_auth_required_when_refresh_fails(cx: &mut TestAppContext) {
689        let client = make_fake_http_client(|_req| {
690            Box::pin(async {
691                Ok(Response::builder()
692                    .status(401)
693                    .header(
694                        "WWW-Authenticate",
695                        r#"Bearer resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource", scope="read write""#,
696                    )
697                    .body(AsyncBody::from(b"Unauthorized".to_vec()))
698                    .unwrap())
699            })
700        });
701
702        // Refresh returns false — no new token available.
703        let provider = FakeTokenProvider::new(Some("stale-token"), false);
704        let transport = HttpTransport::new_with_token_provider(
705            client,
706            "http://mcp.example.com/mcp".to_string(),
707            HashMap::default(),
708            cx.background_executor.clone(),
709            Some(provider.clone()),
710        );
711
712        let err = transport
713            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
714            .await
715            .unwrap_err();
716
717        let transport_err = err
718            .downcast_ref::<TransportError>()
719            .expect("error should be TransportError");
720        match transport_err {
721            TransportError::AuthRequired { www_authenticate } => {
722                assert_eq!(
723                    www_authenticate
724                        .resource_metadata
725                        .as_ref()
726                        .map(|u| u.as_str()),
727                    Some("https://mcp.example.com/.well-known/oauth-protected-resource"),
728                );
729                assert_eq!(
730                    www_authenticate.scope,
731                    Some(vec!["read".to_string(), "write".to_string()]),
732                );
733            }
734        }
735        assert_eq!(provider.refresh_count(), 1);
736    }
737
738    #[gpui::test]
739    async fn test_401_returns_auth_required_without_provider(cx: &mut TestAppContext) {
740        let client = make_fake_http_client(|_req| {
741            Box::pin(async {
742                Ok(Response::builder()
743                    .status(401)
744                    .header("WWW-Authenticate", "Bearer")
745                    .body(AsyncBody::from(b"Unauthorized".to_vec()))
746                    .unwrap())
747            })
748        });
749
750        // No token provider at all.
751        let transport = HttpTransport::new(
752            client,
753            "http://mcp.example.com/mcp".to_string(),
754            HashMap::default(),
755            cx.background_executor.clone(),
756        );
757
758        let err = transport
759            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
760            .await
761            .unwrap_err();
762
763        let transport_err = err
764            .downcast_ref::<TransportError>()
765            .expect("error should be TransportError");
766        match transport_err {
767            TransportError::AuthRequired { www_authenticate } => {
768                assert!(www_authenticate.resource_metadata.is_none());
769                assert!(www_authenticate.scope.is_none());
770            }
771        }
772    }
773
774    #[gpui::test]
775    async fn test_401_after_successful_refresh_still_returns_auth_required(
776        cx: &mut TestAppContext,
777    ) {
778        // Both requests return 401 — the server rejects the refreshed token too.
779        let client = make_fake_http_client(|_req| {
780            Box::pin(async {
781                Ok(Response::builder()
782                    .status(401)
783                    .header("WWW-Authenticate", "Bearer")
784                    .body(AsyncBody::from(b"Unauthorized".to_vec()))
785                    .unwrap())
786            })
787        });
788
789        let provider = FakeTokenProvider::new(Some("token"), true);
790        let transport = HttpTransport::new_with_token_provider(
791            client,
792            "http://mcp.example.com/mcp".to_string(),
793            HashMap::default(),
794            cx.background_executor.clone(),
795            Some(provider.clone()),
796        );
797
798        let err = transport
799            .send(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#.to_string())
800            .await
801            .unwrap_err();
802
803        err.downcast_ref::<TransportError>()
804            .expect("error should be TransportError");
805        // Refresh was attempted exactly once.
806        assert_eq!(provider.refresh_count(), 1);
807    }
808}
809
Served at tenant.openagents/omega Member data and write actions are omitted.