Skip to repository content809 lines · 31.3 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T05:02:23.783Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
http.rs
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