Skip to main content

prism_mcp_rs/transport/
http_auth.rs

1//! HTTP Transport with OAuth 2.1 Authorization Support
2//!
3//! Module extends the HTTP transport with automatic authorization
4//! handling, including token refresh and 401/403 challenge handling.
5
6use async_trait::async_trait;
7use std::sync::Arc;
8
9use crate::auth::{AuthConfig, AuthorizationClient};
10use crate::core::error::{McpError, McpResult};
11use crate::protocol::types::{JsonRpcNotification, JsonRpcRequest, JsonRpcResponse};
12use crate::transport::{http::HttpClientTransport, Transport};
13
14/// HTTP transport with automatic authorization support
15pub struct AuthorizedHttpTransport {
16    /// Base HTTP transport
17    inner: HttpClientTransport,
18    /// Authorization client
19    auth_client: Arc<AuthorizationClient>,
20    /// Whether authorization is enabled
21    auth_enabled: bool,
22}
23
24impl AuthorizedHttpTransport {
25    /// Create a new authorized HTTP transport
26    pub async fn new(
27        base_url: String,
28        sse_url: Option<String>,
29        auth_config: AuthConfig,
30    ) -> McpResult<Self> {
31        let auth_enabled = auth_config.enabled;
32        let inner = HttpClientTransport::new(&base_url, sse_url.as_ref()).await?;
33        let auth_client = Arc::new(AuthorizationClient::new(auth_config, base_url.clone()));
34
35        Ok(Self {
36            inner,
37            auth_client,
38            auth_enabled,
39        })
40    }
41
42    /// Send a request with authorization
43    async fn send_with_auth(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
44        // If auth is disabled, just pass through
45        if !self.auth_enabled {
46            return self.inner.send_request(request).await;
47        }
48
49        // Get or refresh token
50        let token = self.auth_client.get_token().await?;
51
52        // Add authorization header
53        crate::auth::client::add_auth_header(&mut self.inner.headers, &token);
54
55        // Try sending the request
56        match self.inner.send_request(request.clone()).await {
57            Ok(response) => Ok(response),
58            Err(McpError::Http(msg)) if msg.contains("401") || msg.contains("403") => {
59                // Token might be expired, try refreshing
60                let fresh_token = self.auth_client.token_manager().refresh_token().await?;
61
62                // Update header with fresh token
63                crate::auth::client::add_auth_header(&mut self.inner.headers, &fresh_token);
64
65                // Retry the request
66                self.inner.send_request(request).await
67            }
68            Err(e) => Err(e),
69        }
70    }
71
72    /// Handle initial auth challenge response (typically 401 or 403).
73    pub async fn handle_unauthorized(&self, www_authenticate: &str) -> McpResult<String> {
74        self.auth_client.handle_unauthorized(www_authenticate).await
75    }
76
77    /// Handle OAuth callback
78    pub async fn handle_callback(&self, callback_url: &str) -> McpResult<String> {
79        self.auth_client.handle_callback(callback_url).await
80    }
81
82    /// Check if authenticated
83    pub async fn is_authenticated(&self) -> bool {
84        if !self.auth_enabled {
85            return true; // Consider non-auth as authenticated
86        }
87        self.auth_client.is_authenticated().await
88    }
89
90    /// Get the authorization URL to start OAuth flow
91    pub async fn get_authorization_url(&self) -> McpResult<String> {
92        if !self.auth_enabled {
93            return Err(McpError::Auth("Authorization is not enabled".to_string()));
94        }
95
96        // This requires auth server metadata to be already discovered
97        // You might need to trigger discovery first
98        let context = self.auth_client.token_manager().get_context().await;
99
100        if let Some(auth_metadata) = context.auth_server_metadata {
101            self.auth_client
102                .start_authorization_flow(&auth_metadata)
103                .await
104        } else {
105            Err(McpError::Auth(
106                "Authorization server not discovered yet".to_string(),
107            ))
108        }
109    }
110
111    /// Logout and clear tokens
112    pub async fn logout(&self) {
113        self.auth_client.logout().await;
114    }
115}
116
117#[async_trait]
118impl Transport for AuthorizedHttpTransport {
119    async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
120        self.send_with_auth(request).await
121    }
122
123    async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
124        if self.auth_enabled {
125            // Get token and add to headers
126            let token = self.auth_client.get_token().await?;
127            crate::auth::client::add_auth_header(&mut self.inner.headers, &token);
128        }
129
130        self.inner.send_notification(notification).await
131    }
132
133    async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
134        self.inner.receive_notification().await
135    }
136
137    async fn close(&mut self) -> McpResult<()> {
138        self.inner.close().await
139    }
140}
141
142/// Builder for authorized HTTP transport
143pub struct AuthorizedHttpTransportBuilder {
144    base_url: String,
145    sse_url: Option<String>,
146    auth_config: AuthConfig,
147}
148
149impl AuthorizedHttpTransportBuilder {
150    /// Create a new builder
151    pub fn new(base_url: String) -> Self {
152        Self {
153            base_url,
154            sse_url: None,
155            auth_config: AuthConfig::default(),
156        }
157    }
158
159    /// Set SSE URL for notifications
160    pub fn with_sse(mut self, url: String) -> Self {
161        self.sse_url = Some(url);
162        self
163    }
164
165    /// Enable authorization
166    pub fn with_auth(mut self, enabled: bool) -> Self {
167        self.auth_config.enabled = enabled;
168        self
169    }
170
171    /// Set client credentials
172    pub fn with_client_credentials(
173        mut self,
174        client_id: String,
175        client_secret: Option<String>,
176    ) -> Self {
177        self.auth_config.client_id = Some(client_id);
178        self.auth_config.client_secret = client_secret;
179        self
180    }
181
182    /// Set redirect URI
183    pub fn with_redirect_uri(mut self, uri: String) -> Self {
184        self.auth_config.redirect_uri = uri;
185        self
186    }
187
188    /// Set scopes
189    pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
190        self.auth_config.scopes = scopes;
191        self
192    }
193
194    /// Enable dynamic registration
195    pub fn with_dynamic_registration(mut self, enabled: bool) -> Self {
196        self.auth_config.enable_dynamic_registration = enabled;
197        self
198    }
199
200    /// Build the transport
201    pub async fn build(self) -> McpResult<AuthorizedHttpTransport> {
202        AuthorizedHttpTransport::new(self.base_url, self.sse_url, self.auth_config).await
203    }
204}
205
206#[cfg(test)]
207mod tests {
208    use super::*;
209
210    #[tokio::test]
211    async fn test_authorized_transport_creation() {
212        let transport = AuthorizedHttpTransportBuilder::new("https://mcp.example.com".to_string())
213            .with_auth(false) // Disable auth for testing
214            .build()
215            .await;
216
217        assert!(transport.is_ok());
218        let transport = transport.unwrap();
219        assert!(transport.is_authenticated().await);
220    }
221
222    #[tokio::test]
223    async fn test_auth_disabled_passthrough() {
224        let transport = AuthorizedHttpTransportBuilder::new("https://mcp.example.com".to_string())
225            .with_auth(false)
226            .build()
227            .await
228            .unwrap();
229
230        // Should always be authenticated when auth is disabled
231        assert!(transport.is_authenticated().await);
232
233        // Logout should work without error
234        transport.logout().await;
235    }
236}